通俗易读LLM训练-从显存占用分析到DeepSpeed
上个月我蹲在机房,盯着8块A100的显存占用,差点没把键盘砸了。
事情是这样的,我想上手调教一个7B的模型,心里盘算着,这玩意儿显存应该够用吧?结果一算,好家伙,脸都绿了。
你看,我用最笨的方法算一笔账,你听听是不是这个理儿。
咱们就用llama 7B来举例,参数用fp16存,7B乘以2个字节,14G,这把稳了?
不,太天真了。
梯度也得是fp16吧?好,再来14G。
你以为这就完了?最阴险的在后头——那个叫AdamW的优化器,它活像个地主老财,一出手就要占三份fp32地:
一份是模型参数,28G;
一份叫一阶动量m,又是28G;
还有一份叫二阶动量v,还是28G。
你加加看:14 + 14 + 28 + 28 + 28 = 112G?!
我当时就懵了。一张A100满打满算才80G,你一个7B的模型就敢要我112G?这跟让一间屋子住下一整支足球队有什么分别?
说到这儿,你可能会想,那最简单的办法不就是“数据并行”吗?每个GPU都塞一份完整的模型,然后把数据切分不就行了?
对,最朴素,但也最费钱。
你想啊,每个GPU都得有那112G的显存。一张A100都装不下,更别说别的卡了。这感觉就是,你为了让8个人都能看书,每人买了一整套百科全书。又笨重,又浪费,动不动就炸库房。
后来呢?你猜怎么着?DeepSpeed的ZeRO站出来了。
它的核心思想,听起来简直像个笑话,但又妙极了:大家以为是堆硬件,其实是在省冗余。
你看,每张卡都存一份完整的数据,太傻了。咱们分着存,用的时候再拼起来,不就行了?
我花了三天时间,在8卡A100集群上实测了ZeRO的三个阶段,结果简直像看变魔术。
第一招,分割优化器状态。
显存,直接从112G掉到了约40G!怎么做到的?就是把刚才那个地主老财AdamW的fp32家产,分到8张卡上,每张卡只管1/8。通信开销一点没多,跟普通数据并行一模一样。
你看,这就像本来8个人要各自买一套百科全书,现在改成每人只买一本,凑一起就是一套。省了多少?不言而喻。
第二招,再加梯度!
显存直接降到约28G!这回,连反传回来的梯度也分着存了。通信变成了一次“Reduce-Scatter”,再“AllGather”,但数据总量没增加。
这一步最实用!我实测7B模型,只用4张A100就跑起来了!你想想,这对普通人来说,门槛直接砍没了啊!大家以为是高不可攀,其实是触手可及。
第三招,终极魔法——连参数也分了!
显存怎么算? (14 + 14 + 28 + 28 + 28) / 8 = 16G!
16G啊!你想想,从112G到16G,这是差一个数量级的飞跃!但这下代价也来了:通信量多了50%。因为每次前向和反向,都得用AllGather去把参数拼回来。我测8卡的时候,吞吐量比Stage 2低了大概15%。
你看,这就是用通信换显存,一个愿打一个愿挨。
说到这儿,我给你的心里话是这样的:
- **要是你的单卡显存够用,千万别听AI的鬼话上ZeRO,直接数据并行,最稳。**
- **要是显存有点吃紧,但还能忍,Stage 2就是你的绝佳拍档,性价比高得离谱。**
- **要是显存极度紧张,感觉快要爆炸,那就Stage 3,但要做好性能下降的觉悟。**
- **要是你只有一张卡?别怕,还有ZeRO-Offload,把优化器状态扔到CPU内存里,照样能玩。**
我目前用得最多的是Stage 2。4张A100跑7B模型,batch size开到16,稳得像泰山。
最后说一点个人感受,ZeRO这玩意儿,它不是凭空变出显存,而是在展示一个最朴素的道理:我们以为的大多数瓶颈,本质上都是冗余造成的。
你想想,8张卡各存一份完整的优化器状态,就像8个人各建一座图书馆,却只看同一本书。分着存,用的时候拼,这本质上不就是用协作替代重复吗?
不过,也别迷信它。如果模型大到千亿参数,光靠ZeRO还不够,还得上TP/PP这对双雄。
NVIDIA那篇论文里说的,千卡以内ZeRO是王道,千卡以上TP/PP更香。毕竟AllGather的通信开销,就像收过路费的,车多了受不了。
哦对了,差点忘了激活值。前面算显存时,我可没把它算进去。但这玩意儿,才是真正的“隐形杀手”!
多亏了FlashAttention能砍掉那O(S²)的注意力矩阵,还有Ring Attention能把长序列分到多卡上。这些技术跟ZeRO打配合,才能让大模型训练真正跑起来。
训练大模型不是在跟算力抢饭吃——而是在跟自己的认知冗余抢效率。
读者评论 5